Skip to main content

compio_driver\sys\op\socket/
iocp.rs

1use windows_sys::Win32::{
2    Networking::WinSock::{
3        LPFN_ACCEPTEX, LPFN_CONNECTEX, LPFN_DISCONNECTEX, LPFN_GETACCEPTEXSOCKADDRS,
4        LPFN_WSARECVMSG, SO_UPDATE_ACCEPT_CONTEXT, SO_UPDATE_CONNECT_CONTEXT, SOCKADDR,
5        SOCKADDR_STORAGE, SOL_SOCKET, TF_REUSE_SOCKET, WSAID_ACCEPTEX, WSAID_CONNECTEX,
6        WSAID_DISCONNECTEX, WSAID_GETACCEPTEXSOCKADDRS, WSAID_WSARECVMSG, WSAMSG, WSARecv,
7        WSARecvFrom, WSASend, WSASendMsg, WSASendTo, closesocket, setsockopt, socklen_t,
8    },
9    System::IO::OVERLAPPED,
10};
11
12use crate::{OpCode, OpType, sys::op::*};
13
14static ACCEPT_EX: OnceLock<LPFN_ACCEPTEX> = OnceLock::new();
15static GET_ADDRS: OnceLock<LPFN_GETACCEPTEXSOCKADDRS> = OnceLock::new();
16
17const ACCEPT_ADDR_BUFFER_SIZE: usize = std::mem::size_of::<SOCKADDR_STORAGE>() + 16;
18const ACCEPT_BUFFER_SIZE: usize = ACCEPT_ADDR_BUFFER_SIZE * 2;
19
20unsafe impl OpCode for CloseSocket {
21    type Control = ();
22
23    fn op_type(&self, _: &Self::Control) -> OpType {
24        OpType::Blocking
25    }
26
27    unsafe fn operate(&mut self, _: &mut (), _optr: *mut OVERLAPPED) -> Poll<io::Result<usize>> {
28        Poll::Ready(Ok(
29            syscall!(SOCKET, closesocket(self.fd.as_fd().as_raw_fd() as _))? as _,
30        ))
31    }
32}
33
34/// Accept a connection.
35pub struct Accept<S, SA> {
36    pub(crate) fd: S,
37    pub(crate) accept_fd: SA,
38    pub(crate) buffer: [u8; ACCEPT_BUFFER_SIZE],
39}
40
41impl<S, SA> Accept<S, SA> {
42    /// Create [`Accept`]. `accept_fd` should not be bound.
43    pub fn new(fd: S, accept_fd: SA) -> Self {
44        Self {
45            fd,
46            accept_fd,
47            buffer: [0u8; ACCEPT_BUFFER_SIZE],
48        }
49    }
50}
51
52impl<S: AsFd, SA: AsFd> Accept<S, SA> {
53    /// Update accept context.
54    pub fn update_context(&self) -> io::Result<()> {
55        let fd = self.fd.as_fd().as_raw_fd();
56        syscall!(
57            SOCKET,
58            setsockopt(
59                self.accept_fd.as_fd().as_raw_fd() as _,
60                SOL_SOCKET,
61                SO_UPDATE_ACCEPT_CONTEXT,
62                &fd as *const _ as _,
63                std::mem::size_of_val(&fd) as _,
64            )
65        )?;
66        Ok(())
67    }
68
69    /// Get the remote address from the inner buffer.
70    pub fn into_addr(self) -> io::Result<(SA, SockAddr)> {
71        let get_addrs_fn = GET_ADDRS
72            .get_or_try_init(|| {
73                get_wsa_fn(self.fd.as_fd().as_raw_fd(), WSAID_GETACCEPTEXSOCKADDRS)
74            })?
75            .ok_or_else(|| {
76                io::Error::new(
77                    io::ErrorKind::Unsupported,
78                    "cannot retrieve GetAcceptExSockAddrs",
79                )
80            })?;
81        let mut local_addr: *mut SOCKADDR = null_mut();
82        let mut local_addr_len = 0;
83        let mut remote_addr: *mut SOCKADDR = null_mut();
84        let mut remote_addr_len = 0;
85        unsafe {
86            get_addrs_fn(
87                &self.buffer as *const _ as *const _,
88                0,
89                ACCEPT_ADDR_BUFFER_SIZE as _,
90                ACCEPT_ADDR_BUFFER_SIZE as _,
91                &mut local_addr,
92                &mut local_addr_len,
93                &mut remote_addr,
94                &mut remote_addr_len,
95            );
96        }
97        Ok((self.accept_fd, unsafe {
98            SockAddr::new(
99                // SAFETY: the buffer is large enough to hold the address
100                std::mem::transmute::<SOCKADDR_STORAGE, SockAddrStorage>(read_unaligned(
101                    remote_addr.cast::<SOCKADDR_STORAGE>(),
102                )),
103                remote_addr_len,
104            )
105        }))
106    }
107}
108
109unsafe impl<S: AsFd, SA: AsFd> OpCode for Accept<S, SA> {
110    type Control = ();
111
112    unsafe fn operate(&mut self, _: &mut (), optr: *mut OVERLAPPED) -> Poll<io::Result<usize>> {
113        let accept_fn = ACCEPT_EX
114            .get_or_try_init(|| get_wsa_fn(self.fd.as_fd().as_raw_fd(), WSAID_ACCEPTEX))?
115            .ok_or_else(|| {
116                io::Error::new(io::ErrorKind::Unsupported, "cannot retrieve AcceptEx")
117            })?;
118        let mut received = 0;
119        let res = unsafe {
120            accept_fn(
121                self.fd.as_fd().as_raw_fd() as _,
122                self.accept_fd.as_fd().as_raw_fd() as _,
123                self.buffer.sys_slice_mut().ptr() as _,
124                0,
125                ACCEPT_ADDR_BUFFER_SIZE as _,
126                ACCEPT_ADDR_BUFFER_SIZE as _,
127                &mut received,
128                optr,
129            )
130        };
131        win32_result(res, received)
132    }
133
134    fn cancel(&mut self, _: &mut (), optr: *mut OVERLAPPED) -> io::Result<()> {
135        cancel(self.fd.as_fd().as_raw_fd(), optr)
136    }
137}
138
139static CONNECT_EX: OnceLock<LPFN_CONNECTEX> = OnceLock::new();
140
141impl<S: AsFd> Connect<S> {
142    /// Update connect context.
143    pub fn update_context(&self) -> io::Result<()> {
144        syscall!(
145            SOCKET,
146            setsockopt(
147                self.fd.as_fd().as_raw_fd() as _,
148                SOL_SOCKET,
149                SO_UPDATE_CONNECT_CONTEXT,
150                null(),
151                0,
152            )
153        )?;
154        Ok(())
155    }
156}
157
158unsafe impl<S: AsFd> OpCode for Connect<S> {
159    type Control = ();
160
161    unsafe fn operate(&mut self, _: &mut (), optr: *mut OVERLAPPED) -> Poll<io::Result<usize>> {
162        let connect_fn = CONNECT_EX
163            .get_or_try_init(|| get_wsa_fn(self.fd.as_fd().as_raw_fd(), WSAID_CONNECTEX))?
164            .ok_or_else(|| {
165                io::Error::new(io::ErrorKind::Unsupported, "cannot retrieve ConnectEx")
166            })?;
167        let mut sent = 0;
168        let res = unsafe {
169            connect_fn(
170                self.fd.as_fd().as_raw_fd() as _,
171                self.addr.as_ptr().cast(),
172                self.addr.len(),
173                null(),
174                0,
175                &mut sent,
176                optr,
177            )
178        };
179        win32_result(res, sent)
180    }
181
182    fn cancel(&mut self, _: &mut (), optr: *mut OVERLAPPED) -> io::Result<()> {
183        cancel(self.fd.as_fd().as_raw_fd(), optr)
184    }
185}
186
187/// Disconnect a connected socket and reuse it for another connection.
188pub struct Disconnect<S> {
189    pub(crate) fd: S,
190}
191
192impl<S> Disconnect<S> {
193    /// Create [`Disconnect`].
194    pub fn new(fd: S) -> Self {
195        Self { fd }
196    }
197}
198
199static DISCONNECT_EX: OnceLock<LPFN_DISCONNECTEX> = OnceLock::new();
200
201unsafe impl<S: AsFd> OpCode for Disconnect<S> {
202    type Control = ();
203
204    unsafe fn operate(&mut self, _: &mut (), optr: *mut OVERLAPPED) -> Poll<io::Result<usize>> {
205        let disconnect_fn = DISCONNECT_EX
206            .get_or_try_init(|| get_wsa_fn(self.fd.as_fd().as_raw_fd(), WSAID_DISCONNECTEX))?
207            .ok_or_else(|| {
208                io::Error::new(io::ErrorKind::Unsupported, "cannot retrieve DisconnectEx")
209            })?;
210        let res =
211            unsafe { disconnect_fn(self.fd.as_fd().as_raw_fd() as _, optr, TF_REUSE_SOCKET, 0) };
212        win32_result(res, 0)
213    }
214
215    fn cancel(&mut self, _: &mut (), optr: *mut OVERLAPPED) -> io::Result<()> {
216        cancel(self.fd.as_fd().as_raw_fd(), optr)
217    }
218}
219
220#[derive(Default)]
221#[doc(hidden)]
222pub struct RecvControl {
223    pub(crate) slice: SysSlice,
224}
225
226unsafe impl<T: IoBufMut, S: AsFd> OpCode for Recv<T, S> {
227    type Control = RecvControl;
228
229    unsafe fn init(&mut self, ctrl: &mut Self::Control) {
230        ctrl.slice = self.buffer.sys_slice_mut();
231    }
232
233    unsafe fn operate(
234        &mut self,
235        control: &mut Self::Control,
236        optr: *mut OVERLAPPED,
237    ) -> Poll<io::Result<usize>> {
238        let fd = self.fd.as_fd().as_raw_fd();
239        let mut flags = self.flags.bits() as _;
240        let mut received = 0;
241        let res = unsafe {
242            WSARecv(
243                fd as _,
244                &raw const control.slice as _,
245                1,
246                &mut received,
247                &mut flags,
248                optr,
249                None,
250            )
251        };
252        winsock_result(res, received)
253    }
254
255    fn cancel(&mut self, _: &mut Self::Control, optr: *mut OVERLAPPED) -> io::Result<()> {
256        cancel(self.fd.as_fd().as_raw_fd(), optr)
257    }
258}
259
260#[derive(Default)]
261#[doc(hidden)]
262pub struct RecvVectoredControl {
263    pub(crate) slices: Vec<SysSlice>,
264}
265
266unsafe impl<T: IoVectoredBufMut, S: AsFd> OpCode for RecvVectored<T, S> {
267    type Control = RecvVectoredControl;
268
269    unsafe fn init(&mut self, ctrl: &mut Self::Control) {
270        ctrl.slices = self.buffer.sys_slices_mut();
271    }
272
273    unsafe fn operate(
274        &mut self,
275        control: &mut Self::Control,
276        optr: *mut OVERLAPPED,
277    ) -> Poll<io::Result<usize>> {
278        let fd = self.fd.as_fd().as_raw_fd();
279        let mut flags = self.flags.bits() as _;
280        let mut received = 0;
281        let res = unsafe {
282            WSARecv(
283                fd as _,
284                control.slices.as_ptr() as _,
285                control.slices.len() as _,
286                &mut received,
287                &mut flags,
288                optr,
289                None,
290            )
291        };
292        winsock_result(res, received)
293    }
294
295    fn cancel(&mut self, _: &mut Self::Control, optr: *mut OVERLAPPED) -> io::Result<()> {
296        cancel(self.fd.as_fd().as_raw_fd(), optr)
297    }
298}
299
300#[derive(Default)]
301#[doc(hidden)]
302pub struct SendControl {
303    pub(crate) slice: SysSlice,
304}
305
306unsafe impl<T: IoBuf, S: AsFd> OpCode for Send<T, S> {
307    type Control = SendControl;
308
309    unsafe fn init(&mut self, ctrl: &mut Self::Control) {
310        ctrl.slice = self.buffer.sys_slice();
311    }
312
313    unsafe fn operate(
314        &mut self,
315        control: &mut Self::Control,
316        optr: *mut OVERLAPPED,
317    ) -> Poll<io::Result<usize>> {
318        let mut sent = 0;
319        let res = unsafe {
320            WSASend(
321                self.fd.as_fd().as_raw_fd() as _,
322                (&raw const control.slice).cast(),
323                1,
324                &mut sent,
325                self.flags.bits() as _,
326                optr,
327                None,
328            )
329        };
330        winsock_result(res, sent)
331    }
332
333    fn cancel(&mut self, _: &mut Self::Control, optr: *mut OVERLAPPED) -> io::Result<()> {
334        cancel(self.fd.as_fd().as_raw_fd(), optr)
335    }
336}
337
338#[derive(Default)]
339#[doc(hidden)]
340pub struct SendVectoredControl {
341    pub(crate) slices: Vec<SysSlice>,
342}
343
344unsafe impl<T: IoVectoredBuf, S: AsFd> OpCode for SendVectored<T, S> {
345    type Control = SendVectoredControl;
346
347    unsafe fn init(&mut self, ctrl: &mut Self::Control) {
348        ctrl.slices = self.buffer.sys_slices();
349    }
350
351    unsafe fn operate(
352        &mut self,
353        control: &mut Self::Control,
354        optr: *mut OVERLAPPED,
355    ) -> Poll<io::Result<usize>> {
356        let mut sent = 0;
357        let res = unsafe {
358            WSASend(
359                self.fd.as_fd().as_raw_fd() as _,
360                control.slices.as_ptr() as _,
361                control.slices.len() as _,
362                &mut sent,
363                self.flags.bits() as _,
364                optr,
365                None,
366            )
367        };
368        winsock_result(res, sent)
369    }
370
371    fn cancel(&mut self, _: &mut Self::Control, optr: *mut OVERLAPPED) -> io::Result<()> {
372        cancel(self.fd.as_fd().as_raw_fd(), optr)
373    }
374}
375
376#[derive(Default)]
377#[doc(hidden)]
378pub struct RecvFromControl {
379    pub(crate) slice: SysSlice,
380}
381
382unsafe impl<T: IoBufMut, S: AsFd> OpCode for RecvFrom<T, S> {
383    type Control = RecvFromControl;
384
385    unsafe fn init(&mut self, ctrl: &mut Self::Control) {
386        ctrl.slice = self.buffer.sys_slice_mut();
387    }
388
389    unsafe fn operate(
390        &mut self,
391        control: &mut Self::Control,
392        optr: *mut OVERLAPPED,
393    ) -> Poll<io::Result<usize>> {
394        let fd = self.header.fd.as_fd().as_raw_fd();
395        let mut flags = self.header.flags.bits() as _;
396        let mut received = 0;
397        let res = unsafe {
398            WSARecvFrom(
399                fd as _,
400                (&raw const control.slice).cast(),
401                1,
402                &mut received,
403                &mut flags,
404                &raw mut self.header.addr as *mut SOCKADDR,
405                &raw mut self.header.addr_len,
406                optr,
407                None,
408            )
409        };
410        winsock_result(res, received)
411    }
412
413    fn cancel(&mut self, _: &mut Self::Control, optr: *mut OVERLAPPED) -> io::Result<()> {
414        cancel(self.header.fd.as_fd().as_raw_fd(), optr)
415    }
416}
417
418#[derive(Default)]
419#[doc(hidden)]
420pub struct RecvFromVectoredControl {
421    pub(crate) slices: Vec<SysSlice>,
422}
423
424unsafe impl<T: IoVectoredBufMut, S: AsFd> OpCode for RecvFromVectored<T, S> {
425    type Control = RecvFromVectoredControl;
426
427    unsafe fn init(&mut self, ctrl: &mut Self::Control) {
428        ctrl.slices = self.buffer.sys_slices_mut();
429    }
430
431    unsafe fn operate(
432        &mut self,
433        control: &mut Self::Control,
434        optr: *mut OVERLAPPED,
435    ) -> Poll<io::Result<usize>> {
436        let fd = self.header.fd.as_fd().as_raw_fd();
437        let mut flags = self.header.flags.bits() as _;
438        let mut received = 0;
439        let res = unsafe {
440            WSARecvFrom(
441                fd as _,
442                control.slices.as_ptr() as _,
443                control.slices.len() as _,
444                &mut received,
445                &mut flags,
446                &raw mut self.header.addr as *mut SOCKADDR,
447                &raw mut self.header.addr_len,
448                optr,
449                None,
450            )
451        };
452        winsock_result(res, received)
453    }
454
455    fn cancel(&mut self, _: &mut Self::Control, optr: *mut OVERLAPPED) -> io::Result<()> {
456        cancel(self.header.fd.as_fd().as_raw_fd(), optr)
457    }
458}
459
460#[derive(Default)]
461#[doc(hidden)]
462pub struct SendToControl {
463    pub(crate) slice: SysSlice,
464}
465
466unsafe impl<T: IoBuf, S: AsFd> OpCode for SendTo<T, S> {
467    type Control = SendToControl;
468
469    unsafe fn init(&mut self, ctrl: &mut Self::Control) {
470        ctrl.slice = self.buffer.sys_slice();
471    }
472
473    unsafe fn operate(
474        &mut self,
475        control: &mut Self::Control,
476        optr: *mut OVERLAPPED,
477    ) -> Poll<io::Result<usize>> {
478        let mut sent = 0;
479        let res = unsafe {
480            WSASendTo(
481                self.header.fd.as_fd().as_raw_fd() as _,
482                (&raw const control.slice).cast(),
483                1,
484                &mut sent,
485                self.header.flags.bits() as _,
486                self.header.addr.as_ptr().cast(),
487                self.header.addr.len(),
488                optr,
489                None,
490            )
491        };
492        winsock_result(res, sent)
493    }
494
495    fn cancel(&mut self, _: &mut Self::Control, optr: *mut OVERLAPPED) -> io::Result<()> {
496        cancel(self.header.fd.as_fd().as_raw_fd(), optr)
497    }
498}
499
500#[derive(Default)]
501#[doc(hidden)]
502pub struct SendToVectoredControl {
503    pub(crate) slices: Vec<SysSlice>,
504}
505
506unsafe impl<T: IoVectoredBuf, S: AsFd> OpCode for SendToVectored<T, S> {
507    type Control = SendToVectoredControl;
508
509    unsafe fn init(&mut self, ctrl: &mut Self::Control) {
510        ctrl.slices = self.buffer.sys_slices();
511    }
512
513    unsafe fn operate(
514        &mut self,
515        control: &mut Self::Control,
516        optr: *mut OVERLAPPED,
517    ) -> Poll<io::Result<usize>> {
518        let mut sent = 0;
519        let res = unsafe {
520            WSASendTo(
521                self.header.fd.as_fd().as_raw_fd() as _,
522                control.slices.as_ptr() as _,
523                control.slices.len() as _,
524                &mut sent,
525                self.header.flags.bits() as _,
526                self.header.addr.as_ptr().cast(),
527                self.header.addr.len(),
528                optr,
529                None,
530            )
531        };
532        winsock_result(res, sent)
533    }
534
535    fn cancel(&mut self, _: &mut Self::Control, optr: *mut OVERLAPPED) -> io::Result<()> {
536        cancel(self.header.fd.as_fd().as_raw_fd(), optr)
537    }
538}
539
540static WSA_RECVMSG: OnceLock<LPFN_WSARECVMSG> = OnceLock::new();
541
542#[derive(Default)]
543#[doc(hidden)]
544pub struct RecvMsgControl {
545    msg: WSAMSG,
546    #[allow(dead_code)]
547    slices: Vec<SysSlice>,
548}
549
550unsafe impl<T: IoVectoredBufMut, C: IoBufMut, S: AsFd> OpCode for RecvMsg<T, C, S> {
551    type Control = RecvMsgControl;
552
553    unsafe fn init(&mut self, ctrl: &mut Self::Control) {
554        ctrl.slices = self.buffer.sys_slices_mut();
555        ctrl.msg.dwFlags = self.header.flags.bits() as _;
556        ctrl.msg.name = &raw mut self.header.addr as _;
557        ctrl.msg.namelen = self.header.addr.size_of() as _;
558        ctrl.msg.lpBuffers = ctrl.slices.as_mut_ptr() as _;
559        ctrl.msg.dwBufferCount = ctrl.slices.len() as _;
560        ctrl.msg.Control = self.control.sys_slice_mut().into_inner();
561    }
562
563    unsafe fn operate(
564        &mut self,
565        control: &mut RecvMsgControl,
566        optr: *mut OVERLAPPED,
567    ) -> Poll<io::Result<usize>> {
568        let recvmsg_fn = WSA_RECVMSG
569            .get_or_try_init(|| get_wsa_fn(self.header.fd.as_fd().as_raw_fd(), WSAID_WSARECVMSG))?
570            .ok_or_else(|| {
571                io::Error::new(io::ErrorKind::Unsupported, "cannot retrieve WSARecvMsg")
572            })?;
573
574        let mut received = 0;
575        let res = unsafe {
576            recvmsg_fn(
577                self.header.fd.as_fd().as_raw_fd() as _,
578                &raw mut control.msg,
579                &mut received,
580                optr,
581                None,
582            )
583        };
584        winsock_result(res, received)
585    }
586
587    fn cancel(&mut self, _: &mut Self::Control, optr: *mut OVERLAPPED) -> io::Result<()> {
588        cancel(self.header.fd.as_fd().as_raw_fd(), optr)
589    }
590
591    unsafe fn set_result(
592        &mut self,
593        control: &mut Self::Control,
594        _: &io::Result<usize>,
595        _: &crate::Extra,
596    ) {
597        self.return_flags = ReturnFlags::from_bits_retain(control.msg.dwFlags);
598        self.header.addr_len = control.msg.namelen as socklen_t;
599        self.control_len = control.msg.Control.len as _;
600    }
601}
602
603#[derive(Default)]
604#[doc(hidden)]
605pub struct SendMsgControl {
606    msg: WSAMSG,
607    #[allow(dead_code)]
608    slices: Vec<SysSlice>,
609}
610
611unsafe impl<T: IoVectoredBuf, C: IoBuf, S: AsFd> OpCode for SendMsg<T, C, S> {
612    type Control = SendMsgControl;
613
614    unsafe fn init(&mut self, ctrl: &mut Self::Control) {
615        ctrl.slices = self.buffer.sys_slices();
616        let control = if self.control.buf_len() == 0 {
617            SysSlice::null()
618        } else {
619            self.control.sys_slice()
620        };
621
622        ctrl.msg.lpBuffers = ctrl.slices.as_ptr() as _;
623        ctrl.msg.dwBufferCount = ctrl.slices.len() as _;
624        ctrl.msg.Control = control.into_inner();
625        if let Some(addr) = &self.addr {
626            ctrl.msg.name = addr.as_ptr() as _;
627            ctrl.msg.namelen = addr.len() as _;
628        }
629    }
630
631    unsafe fn operate(
632        &mut self,
633        control: &mut Self::Control,
634        optr: *mut OVERLAPPED,
635    ) -> Poll<io::Result<usize>> {
636        let mut sent = 0;
637        let res = unsafe {
638            WSASendMsg(
639                self.fd.as_fd().as_raw_fd() as _,
640                &raw mut control.msg,
641                self.flags.bits() as _,
642                &mut sent,
643                optr,
644                None,
645            )
646        };
647        winsock_result(res, sent)
648    }
649
650    fn cancel(&mut self, _: &mut Self::Control, optr: *mut OVERLAPPED) -> io::Result<()> {
651        cancel(self.fd.as_fd().as_raw_fd(), optr)
652    }
653}